Repository navigation
[Feat]: support dspark - #4951
RunningLeon wants to merge 11 commits into
Conversation
245beef to
81b55dc
Compare
81b55dc to
fc9aae4
Compare
Require an explicit host-known query extent instead of reading a CUDA scalar for arange. Use packed input size for ragged/chunked prefill and preserve the padded rectangular graph capacity. Add metadata mapping, extent contract, synchronization guard and CUDA Graph replay regression tests.
Coordinate block-query participation and padding across DP ranks, isolate query metadata, and keep context KV materialization local. Preserve block-zero dummy cache writes, reserve method-specific KV lookahead, and fix V4 attention/MoE parallel layouts. Include config compatibility and regression coverage; keep the independent TP sleep fix out of this change. Validation: 493 CPU tests passed, 54 skipped. Full DP4/EP4 CUDA Graph GSM8K completed for Qwen3.5 and DeepSeek-V4 DSpark.
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
Critical state-cache issues and moderate correctness/test issues remain unresolved.
Get a fresh assessment by requesting another Copilot review.
Review effort: Lite
Findings: 2
Open (4)
What changed in this PR
Adds DSpark speculative decoding with DeepSeek-V4 integration, distributed execution, CUDA graph support, and kernel/cache updates.
Changes:
- Adds DSpark/DFlash configuration, proposers, models, heads, registration, CLI, and documentation.
- Extends V4 state caches, attention, compression, MoE, CUDA graphs, and distributed execution.
- Adds broad tests for speculative decoding, kernels, engines, and model configuration.
| File | Summary | Findings |
|---|---|---|
tests/pytorch/spec_decode/test_strategies.py |
Tests DSpark allocation sizing. | — |
tests/pytorch/spec_decode/test_spec_agent.py |
Tests proposal resolution and warmup metadata. | — |
tests/pytorch/spec_decode/test_reject_sample.py |
Tests deterministic draft sampling. | — |
tests/pytorch/spec_decode/test_dflash_utils.py |
Tests DFlash/DSpark configuration. | — |
tests/pytorch/spec_decode/test_cudagraph_strategy.py |
Tests rectangular graph sizing. | — |
tests/pytorch/nn/test_v4_prefill_metadata.py |
Tests V4 prefill metadata. | — |
tests/pytorch/nn/test_moe_deepep.py |
Tests distributed MoE gathering. | — |
tests/pytorch/kernel/test_v4_flatten_kv.py |
Tests expanded V4 rings. | — |
tests/pytorch/kernel/test_v4_compressor.py |
Tests mixed compression windows. | — |
tests/pytorch/kernel/test_gated_delta_rule_split_v.py |
Tests split-V recurrent kernels. | Nit (2 votes): cases do not exercise the num_states == 8 split-V specialization. |
tests/pytorch/kernel/test_causal_conv1d_strided.py |
Tests strided convolution updates. | Moderate (1 vote): fake output allocation is non-contiguous and fails before exercising the implementation. |
tests/pytorch/engine/test_model_agent.py |
Tests DSpark weights and state transactions. | — |
tests/pytorch/engine/test_model_agent_kv_connector.py |
Updates model-agent fixtures. | — |
tests/pytorch/engine/test_executor_base.py |
Tests state-cache synchronization. | — |
tests/pytorch/engine/test_cache_engine/test_state_engine.py |
Tests state rollback. | — |
tests/pytorch/config/test_model_config.py |
Tests V4 ring geometry. | — |
lmdeploy/pytorch/strategies/ar_spec/engine.py |
Supports DSpark query allocation. | — |
lmdeploy/pytorch/strategies/ar_spec/cudagraph.py |
Sizes rectangular speculative graphs. | — |
lmdeploy/pytorch/strategies/ar_spec/__init__.py |
Builds DSpark strategies. | — |
lmdeploy/pytorch/spec_decode/spec_agent.py |
Adds draft state-cache support. | — |
lmdeploy/pytorch/spec_decode/proposers/dspark.py |
Implements DSpark proposal execution. | — |
lmdeploy/pytorch/spec_decode/proposers/dflash.py |
Adds block-parallel DFlash execution. | — |
lmdeploy/pytorch/spec_decode/proposers/base.py |
Propagates named state caches. | — |
lmdeploy/pytorch/spec_decode/proposers/__init__.py |
Registers DSpark. | — |
lmdeploy/pytorch/spec_decode/dspark_utils.py |
Parses and validates DSpark configuration. | — |
lmdeploy/pytorch/spec_decode/dflash_utils.py |
Updates distributed validation. | — |
lmdeploy/pytorch/spec_decode/block_parallel.py |
Adds DP block-query padding. | — |
lmdeploy/pytorch/spec_decode/base.py |
Resolves proposal protocols. | — |
lmdeploy/pytorch/nn/v4_compressor.py |
Passes state IDs to compression. | — |
lmdeploy/pytorch/nn/v4_attention.py |
Exposes V4 cache-write operations. | — |
lmdeploy/pytorch/nn/moe/v4_fp4.py |
Gathers DP rows before TP experts. | — |
lmdeploy/pytorch/models/utils/cudagraph.py |
Stores graph KV bounds. | — |
lmdeploy/pytorch/models/qwen3_dspark.py |
Adds the Qwen DSpark model. | — |
lmdeploy/pytorch/models/qwen3_dflash.py |
Handles nullable target hidden size. | — |
lmdeploy/pytorch/models/module_map.py |
Registers DSpark architectures. | — |
lmdeploy/pytorch/models/glm_moe_dsa.py |
Captures auxiliary hidden states. | — |
lmdeploy/pytorch/models/dspark_heads.py |
Implements DSpark Markov heads. | — |
lmdeploy/pytorch/models/deepseek_v4.py |
Adds V4 auxiliary states and ring capacity. | — |
lmdeploy/pytorch/models/deepseek_v4_dspark.py |
Adds bundled V4 DSpark stages. | — |
lmdeploy/pytorch/models/deepseek_v32.py |
Adds auxiliary-state capture. | — |
lmdeploy/pytorch/model_inputs.py |
Adds block-plan metadata. | — |
lmdeploy/pytorch/kernels/cuda/v4_flatten_kv.py |
Separates logical and physical ring sizes. | — |
lmdeploy/pytorch/kernels/cuda/v4_compressor.py |
Handles mixed windows and padding. | — |
lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py |
Adds split-V specialization. | — |
lmdeploy/pytorch/kernels/cuda/causal_conv1d.py |
Supports strided layouts. | — |
lmdeploy/pytorch/engine/model_agent/dp_utils.py |
Serializes block-query readiness. | — |
lmdeploy/pytorch/engine/model_agent/agent.py |
Adds DSpark transactions and DP planning. | — |
lmdeploy/pytorch/engine/executor/base.py |
Reserves draft state-cache memory. | — |
lmdeploy/pytorch/engine/config_builder.py |
Configures DSpark distribution constraints. | — |
lmdeploy/pytorch/engine/cache_engine/state.py |
Implements V4 state snapshots and rollback. | Critical (1 vote): padded DP/EP queries can use state_offsets == -1 and alias a live state slot. |
lmdeploy/pytorch/consts.py |
Defines ring alignment. | — |
lmdeploy/pytorch/configurations/qwen3_5.py |
Allows DSpark configuration. | — |
lmdeploy/pytorch/configurations/deepseek_v4.py |
Builds V4 ring geometry. | — |
lmdeploy/pytorch/config.py |
Adds DSpark and V4 fields. | — |
lmdeploy/pytorch/backends/cuda/v4_compressor.py |
Passes state IDs to CUDA compression. | — |
lmdeploy/pytorch/backends/cuda/op_backend.py |
Propagates ring capacity. | — |
lmdeploy/pytorch/backends/cuda/graph_runner/runner.py |
Adds graph KV capacity. | — |
lmdeploy/pytorch/backends/cuda/graph_runner/full_graph.py |
Restores state caches during capture. | Critical (1 vote): snapshot indexing assumes dimension 1, which is invalid for anonymous SSM caches. |
lmdeploy/pytorch/backends/cuda/causal_conv1d.py |
Enables strided TileLang updates. | Moderate (1 vote): empty_like can return a non-contiguous output, violating the decode contract. |
lmdeploy/pytorch/backends/cuda/attention/v4.py |
Adds rectangular V4 decode execution. | — |
lmdeploy/pytorch/backends/cuda/attention/v4_utils.py |
Makes prefill extents host-known. | — |
lmdeploy/pytorch/backends/compressor.py |
Extends the compressor interface. | — |
lmdeploy/pytorch/backends/attention.py |
Extends V4 metadata and cache interfaces. | — |
lmdeploy/cli/utils.py |
Exposes the DSpark CLI option. | — |
docs/zh_cn/advance/spec_decoding.md |
Documents DSpark usage. | Nit (1 vote): update DP/EP limitations to match supported configurations. |
docs/en/advance/spec_decoding.md |
Documents DSpark usage. | Nit (3 votes): update DP/EP limitations to match supported configurations. |
benchmark/profile_restful_api.py |
Corrects streamed completion lengths. | — |
.gitignore |
Ignores DeepGEMM artifacts. | — |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
| named_state_caches = getattr(context, 'named_state_caches', None) | ||
| state_ids = self.meta.input_buffers.get('state_ids') | ||
| if named_state_caches and state_ids is not None: | ||
| active_state_ids = torch.unique(state_ids[state_ids >= 0]) | ||
| for cache in named_state_caches.values(): | ||
| state_backups.append( | ||
| (cache, cache.index_select(1, active_state_ids).clone())) |
| f'{max_q_seqlen}, capacity {capacity}.') | ||
|
|
||
| # Layer-scoped V4 state views use [layers, slots, rows, ...]. | ||
| before = cache[:, state_offsets[:, None], rows].clone() |
| debugging fallback. DSpark V1 requires `dp=1` and `ep=1`. Prefix caching, draft | ||
| KV-cache quantization, guided decoding, and confidence-based dynamic | ||
| verification are not supported in this fixed-window version. The |
| states = batch + 2 | ||
| initial = torch.randn(states, 8, 32, 128, 128, device='cuda') * .01 | ||
| backing = torch.empty(states, 3, 8, 32, 128, 128, device='cuda') |
| attn_metadata.kv_seqlens = input_buffers['kv_seqlens'] | ||
| attn_metadata.cu_seqlens_q = input_buffers['cu_seqlens_q'] | ||
| attn_metadata.cu_seqlens_k = input_buffers['cu_seqlens_k'] | ||
| attn_metadata.is_cuda_graph = True |
There was a problem hiding this comment.
Do we really need these 3 fields?
max_kv_seqlen should be equal to block_offsets.size(1) * block_size
graph_sum_kv_seqlen has not been used.
| if has_state_ids: | ||
| state_id = tl.load(state_ids_ptr + batch_id) | ||
| if state_id < 0: | ||
| return |
There was a problem hiding this comment.
I am not sure if it is worth to pass state_id just for skip padding request.
| v_per_warp = T.ceildiv(min(desired, max_v_per_warp), min_v_per_warp) * min_v_per_warp | ||
| v_per_warp = max(v_per_warp, min_v_per_warp) | ||
| target_v_per_cta = V | ||
| target_v_per_cta = V if v_tile_size is None else min(V, v_tile_size) |
There was a problem hiding this comment.
Any benchmark result about this optimization.
| dist_config: DistConfig = field(default_factory=DistConfig) | ||
| target_layer_ids: tuple[int, ...] | None = None | ||
| mask_token_id: int | None = None | ||
| dspark_sample_from_anchor: bool | None = None |
There was a problem hiding this comment.
I think we can make a sub config in SpecDecodeConfig for DSpark.
| self.scheduler_config = scheduler_config | ||
| self.cache_config = cache_config | ||
| self.num_spec_tokens = num_spec_tokens | ||
| self.draft_query_len = draft_query_len |
There was a problem hiding this comment.
We can rename it to draft_kv_lookahead, initialize it to num_spec_tokens if it is None. Then we can simplify the design of get_num_required_tokens and get_prealloc_size.
| ) | ||
| validate_dflash_dist_config(dist_config) | ||
| validate_dflash_runtime_config(cache_config=cache_config, backend_config=engine_config) | ||
| elif speculative_config.method in ('dflash', 'dspark'): |
There was a problem hiding this comment.
As you will import validate function from spec_decode, I think you can warp all branch about spec config into spec_decode so builder do not need to know the detail about speculative_config.
| named_state_caches = getattr(context, 'named_state_caches', None) | ||
| state_ids = self.meta.input_buffers.get('state_ids') | ||
| if named_state_caches and state_ids is not None: | ||
| active_state_ids = torch.unique(state_ids[state_ids >= 0]) |
There was a problem hiding this comment.
unique and masked select would lead to cuda synchronize.


Motivation
This PR has two goals: add DSpark as a first-class speculative decoding method in the PyTorch engine, and enable data/expert parallel execution for both DFlash and DSpark. DSpark supports external Speculators-format draft checkpoints for Qwen-family targets and DSpark weights bundled with DeepSeek-V4 checkpoints. Both block speculative methods retain CUDA Graph execution in distributed serving.
Modification
mtp.*DSpark weights.Current limitations:
Use cases
External DSpark checkpoint:
Bundled DeepSeek-V4 DSpark weights:
Checklist
git diff --checkpass.